Skip to content

Fixing several bugs in Halide's AMX support - #8350

Draft
yihozhang wants to merge 14 commits into
halide:mainfrom
yihozhang:yihozhang-amx-extraction-bugs
Draft

yihozhang wants to merge 14 commits into
halide:mainfrom
yihozhang:yihozhang-amx-extraction-bugs

Conversation

@yihozhang

@yihozhang yihozhang commented Jul 22, 2024 •

Copy link
Copy Markdown

This PR exposes and fixes a number of issues with Halide's current AMX support:

  • Ensures that the modulo in RHS matrix is correct (tiled_matmul_wrong_modulo.cpp)
  • Removes the stride in the 1D case, since there will be just one row and the stride does not really matter.
  • Checks that the RHS tile to be loaded is contiguous in its rows (tiled_matmul_wrong_layout.cpp)
  • Ensures that the tiling of RHS is 4 / element_width (tiled_matmul_wrong_tiling.cpp, tiled_matmul_wrong_pattern.cpp)
  • Swaps tile_x and tile_y in the 3D case
    • this was not found in the first place since tile_x=tile_y always holds in the original tests

Not yet addressed

  • The second child to the Add node in the matcher is ignored.
  • support cases where tile_x = 1 ortile_y=1
  • Load's predicate is not checked
  • Make clear that the optional Broadcast pattern should only occur for the 1D and 2D case

@yihozhang yihozhang changed the title Bugs in Halide's AMX support Fixing several bugs in Halide's AMX support Aug 13, 2024
abadams added a commit that referenced this pull request Jul 13, 2026
* Specialized x86 implementation of interleave_vectors

* Update test to be more exhaustive

* Fix comment.

The previous comment reported a time that seemed to have regressed. It
was not 8.2ms on main - more like 11

* Comment fix

* clang-tidy fixes

* Make variable names more consistent

* Simplify code with helper lambda

* Comment tweaks

* Don't do half-width unpcks

* Use optimization fences in the base class too

Before:

Computing best tile sizes for each type
.................................................
bytes, tile width, tile height, bandwidth (GB/s):
1 8 8 20.9997
1 16 8 20.8329
1 8 16 18.5702
1 8 32 17.2463
1 8 64 14.312

2 8 16 19.2047
2 8 8 18.8368
2 16 8 17.0593
2 8 32 17.0591
2 4 8 15.7681

4 8 8 24.9364
4 4 16 22.9699
4 8 16 22.5743
4 4 32 22.255
4 4 8 20.4468

8 8 8 38.4094
8 16 4 28.4167
8 16 8 27.6184
8 8 4 27.6062
8 8 16 26.8693

After:

Computing best tile sizes for each type
.................................................
bytes, tile width, tile height, bandwidth (GB/s):
1 16 32 34.1921
1 16 16 31.8399
1 8 16 25.575
1 16 64 25.1665
1 32 16 25.0061

2 8 32 28.2635
2 8 16 27.7648
2 16 16 27.2126
2 16 32 23.9034
2 8 8 23.6345

4 8 16 34.5303
4 8 8 28.3653
4 16 8 26.8521
4 8 32 26.084
4 16 16 24.4519

8 8 8 33.7163
8 8 4 29.1339
8 4 16 26.418
8 16 4 25.4663
8 2 8 24.3949

* Use Catanzaro's algorithm for non-power-of-two interleaves

* Support more interleave and deinterleave patterns

* clang-tidy fix

* Handle multiple let injections at same site

Also better algorithm for innermost containing stmt

* better simplification and better handling of composite factors

* Fix innermost_containing_node

* Fix some simd op check failures

* Fix infinite recursion issue and missed case in interleave codegen

* Adjust expectations in stage_strided_loads test

* Allow reversed suffix or not in sve test

* Don't use optimization fences on hexagon

* Fix infinite simplifier loop

* Don't hoist transposes on hexagon

* Make distinct strided load nodes in the IR distinct in memory too

* arm-32 has no vst2 for 64-bit elements

* Windows bad filename fix in simd op check

* Temporary dumping of cpu info to debug github actions issue

* dump cpuinfo in makefile testing workflow

To help diagnose occasional illegal instruction errors

* Address review comments

* Remove duplicate function body

* Use slice of predicate

* clang-format

* SVE fixes

Co-authored-by: Claude Code <noreply@anthropic.com>

* Move optimization_fence back

* Try to thread the needle with webassembly nonsense

* Fix msvc warning

* Skip simd_op_check_sve2 on old llvms

* Skip test on sve2 with llvm 21

* Skip block transpose performance test for sve2 on llvm 21

* Skip sub-test that triggers llvm bug

* Test should hopefully now work with llvm main

* Introduce MultiRamp, a multi-dimensional ramp abstraction

Adds a MultiRamp IR helper that generalizes the old InterleavedRamp: a
nested ramp with a scalar base, a vector of strides (innermost first),
and a vector of per-dim lane counts. Supports in-place mul/add/div/mod
with symbolic strides where possible, reorder, slice, flatten into 1D
ramps, shuffle index construction for permutations and slices, and an
alias-free predicate.

Replaces InterleavedRamp recognition and handling in VectorizeLoops with
MultiRamp. The reduction-store path peels stride-zero and non-alias-free
dims (turning the latter into unrolled containing loops), computes the
per-iteration shuffle mask from the pre-peel shape via
shuffle_from_slice, and gracefully falls back when alias-freedom can't
be proven.

Wires MultiRamp into the simplifier: the Load and Store rules that
recognize a ramp-of-ramp index now use MultiRamp to rotate the stride-1
dim outermost via a single Shuffle::make_transpose, fixing a latent
correctness bug in the old rule for triply-nested ramps.

In FlattenNestedRamps, teaches the Load and Store visitors to recognize
multiramp indices and emit a concat of per-outer-multi-index 1D ramp
loads/stores, rather than a single flat-indexed load that downstream
passes struggle to combine. The existing bounded-span-to-dense-load
path runs first so strided-gather patterns (e.g. HVX vdelta) are
preserved.

Adds correctness tests for the MultiRamp API (test/correctness/multiramp.cpp)
and a nested-vectorize reduction test (transposed_vector_reduce.cpp).

Co-authored-by: Claude <noreply@anthropic.com>

* Clarify MultiRamp API and simplify the atomic-store reduction path

API and doc cleanup on MultiRamp:
 - Add a real invariants block to the class header; clarify that 0-dim
   (scalar) multiramps are legal and all methods handle them.
 - Switch all class member comments to doxygen style.
 - Reword alias_free, alias_free_slice, rotate_stride_one_innermost, and
   the shuffle_from_* overloads so each leads with what it does.
 - Fix is_multiramp's Mul branch to use fresh local MultiRamps rather
   than letting a failed first attempt leak partial state into the
   second.
 - Drop the single-dim shuffle_from_slice overload in favour of the
   multi-dim version.
 - Relax add() so it handles 0-dim inputs trivially (base+base); this
   also makes operator== work for 0-dim via its existing add path.
 - Add accept/mutate methods (Function-idiom) so callers don't reach
   into base/strides to walk scalar subexpressions.
 - Add alias_free_slice (replaces ad-hoc in-caller peeling) and
   rotate_stride_one_innermost (replaces near-duplicate dance in the
   simplifier rules).

Use those APIs from VectorizeLoops, FlattenNestedRamps, Simplify_Exprs,
and Simplify_Stmts. In particular the atomic-store reduction block in
VectorizeLoops is restructured: one alias_free_slice call discovers
both stride-zero peels (handled via VectorReduce or a tree reduction
over a reordered b) and symbolic/overlapping aliasing peels (handled
via an unrolled cartesian-product loop block). The b's current lane
layout is tracked as a MultiRamp (b_shape_mr) so the per-iteration
slice of the reduced vector can be computed via shuffle_from_slice.

Extend the downsampling/atomic-vectorize test (downsampling_reduce.cpp)
to exercise MultiRamp::div through the vectorize path, and expand the
multiramp API tests.

Co-authored-by: Claude <noreply@anthropic.com>

* Add for_each_coordinate helper and use it

Three places in Halide iterate the cartesian product of a box of
integer sizes — MultiRamp::shuffle_from_permuted, MultiRamp::flatten,
MultiRamp::shuffle_from_slice, and the unroll block in VectorizeLoops's
atomic-store reduction path. Replace each manual decompose-via-
rem%/rem/ loop with a call to a new for_each_coordinate helper in
Util.h that invokes a callback on each coordinate in lex order.

Co-authored-by: Claude <noreply@anthropic.com>

* Fix three bugs found by the randomized vectorized-reduction test

is_multiramp's recursive impl may leave its output in a partial state
on failure (e.g. after a successful Ramp::base recursion followed by a
failed stride check). The contract previously told callers "don't read
*result on failure," but honoring that required every recursive call
site to use a fresh local MultiRamp. Split is_multiramp into an
internal impl that keeps the old partial-state behavior plus a thin
public wrapper that commits to *result only on success. The public
contract is now the cleaner "untouched on failure." The Mul branch's
fresh-local workaround falls out.

VectorizeLoops's atomic-store reduction path needs to shuffle `b` from
its original lane order into a permuted order (so a subsequent
reduction tree can slice contiguous sub-vectors per stride-zero peel).
It was calling shuffle_from_permuted directly, but that method returns
indices for the opposite direction — "permuted → original," which is
what the Simplify_Exprs caller wants. Invert the result as a
permutation. Without this, any case with multiple stride-zero peels
that needed a real reorder summed lanes that should have gone to
different output addresses.

Simplify_Shuffle's "slice of concat" rule drops concat vectors that
don't overlap with the slice's range. The overlap check tested whether
the concat vector's start OR its last lane was inside the slice, but
missed the case where the slice is entirely contained within one
concat vector (neither endpoint inside). When every vector contained
the slice, new_concat_vectors came out empty and Shuffle::make_concat
tripped its empty-vector assert. Replaced with a standard interval-
overlap check. This is a pre-existing bug that the randomized test
exercised through multiramp-derived slice patterns.

Co-authored-by: Claude <noreply@anthropic.com>

* Expand MultiRamp and vectorized-reduction tests

Add hand-picked tests for MultiRamp API properties that weren't
previously covered: mul, operator==, alias_free_slice (unique lanes /
zero-stride peeling / degenerate scalar), rotate_stride_one_innermost
(rotation + transpose round-trip), and is_multiramp round-trips for a
handful of shapes.

Add test_random to transposed_vector_reduce.cpp: 1000 random
quasi-affine store/load index pairs over a 3-dim RDom, each compiled
scalarly and with .atomic().vectorize() across all three RVars and
compared. This test found all three bugs fixed in the preceding
commit.

Co-authored-by: Claude <noreply@anthropic.com>

* Cut down on number of transposed_vector_reduce test cases

* Clarify comments flagged by a weak-model review

Ran a weak subagent (Haiku) over the MultiRamp PR as an adversarial
comprehension test — asking it to explain the code in detail, then
fixing whatever it got wrong. The theory: if a weaker model
misreads something, the comment is probably unclear, not the model.

Fixes prompted by the review:

- Simplify_Exprs.cpp / Simplify_Stmts.cpp: stale "outermost" wording
  from before rotate_stride_one_outermost was renamed to
  rotate_stride_one_innermost. The comments contradicted the function
  name and Haiku echoed the contradiction.
- MultiRamp.h alias_free: state explicitly that the returned Expr is
  a sufficient (not necessary) condition for lane uniqueness.
- MultiRamp.h alias_free_slice: clarify that kept dims are a subset
  preserving order, not necessarily a prefix.
- VectorizeLoops.cpp: rename ContainingLoop -> UnrolledLoop and note
  that the peeled dims are fully unrolled into a flat Block, not a
  runtime loop nest (despite the old name).
- MultiRamp.h alias_free_slice: note that stride-zero and purely
  symbolic strides always peel (added by Andrew directly).

A second Haiku pass after these edits answered every question
correctly, including the ones it got wrong the first time.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>

* Remove holdover from transpose branch

* Apply pre-commit auto-fixes

* Fixes for the pre-commit auto fixes

* clang-tidy fix

* Skip test under SVE

* This level of staging now unnecessary

* Simplify new simplifier rules

* Rewrite ExtractTileOperations to use MultiRamp

Replaces the four hand-rolled AST pattern matchers in
ExtractTileOperations with a uniform multiramp-based approach. The
previous matchers were brittle (Add operand ordering, tile_x/tile_y
swap, missing stride checks, etc. — see #8350) and had separate
branches for the broadcast / collapsed / general RHS cases.

The new flow lifts each load index to a MultiRamp, then coerces it
into the canonical AMX shape via MultiRamp::strides_for_shape (a new
one-sided gcd-walk that returns per-target-dim strides), and reads
off the extracted strides from known slot positions. The inner-K /
broadcast / row-stride bits all fall out of the same path.

Also adds:
- A normalizing MultiRamp constructor that drops extent-1 dims.
- MultiRamp::strides_for_shape to map an MR's lane sequence onto a
  caller-specified dim shape.
- Asymmetric tile-size cases to the correctness test (tile_x !=
  tile_y) — these would have caught the swap bug listed in #8350.
- Constraints on the perf-test ImageParams' VNNI inner stride so the
  contiguity check now folds away cleanly.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>

* typo fix

* Move and dedup success prints

* Remove bad assert (e may be undefined)

* Fully relower intrinsics to help multiramp matching

* Support outer tiling

* Better error handling

* Add more test cases and fix bugs found

* clang-tidy fix

* Revert unintended changes

* Better error message if no update def.

* Add error-path test suite for ExtractTileOperations

A correctness test that exercises ten of the user-facing error paths in
ExtractTileOperations.cpp. Each scenario is the most natural-looking
matmul pattern that triggers a particular reject, doubling as a TODO
list of cases we'd ideally support but currently don't:

- too_large            tile_x > 16
- bad_result_type      i8 * i8 -> i16 (AMX always accumulates i32/f32)
- naive_rhs            row-major RHS without VNNI packing
- indirect             gather-style A(r, row_indices(y)) * B(...)
- conv1d               1D conv of a 2D signal (LHS depends on x, k, y)
- no_matmul            store_in(AMXTile) on a non-matmul Func
- widening_16bit       i16 * i16 -> i32 (only 8-bit inputs supported)
- inconsistent_tiles   one Func with two updates at different tile sizes
- not_a_matmul_pattern row-sum into AMXTile (no multiply)
- scaled_matmul        A(r, y) * 3 (RHS hoisted out of the reduce)

The harness wraps each scenario in try/catch around Halide::CompileError.
Halide is sometimes built without exceptions; in that case the test
prints [SKIP] and exits 0 since the catch path can't fire.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>

* tidy fixes

* Fix use-after-free

* Address PR review TODOs

1. Add is_tile_memory_type(MemoryType) helper, and use it at the four
   sites in StageStridedLoads and Deinterleave that previously
   hard-coded `mt == AMXTile` checks to opt out of load/store rewrites.
   Currently the helper just returns `t == AMXTile`, but future
   natively-2D memory types (e.g. other vendors' matrix accumulators)
   can be added by extending it in one place.

2. Guard tiled_matmul_errors with #if HALIDE_WITH_EXCEPTIONS so the
   compilation doesn't fail when exceptions are disabled — matching
   the pattern used by other exception-using correctness tests
   (exception.cpp, bad_partition_always_throws.cpp, etc.). The runtime
   exceptions_enabled() check is still there so the test also handles
   the rare case of being built with HALIDE_WITH_EXCEPTIONS but linked
   against a libHalide that wasn't.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>

---------

Co-authored-by: Martijn Courteaux <courteauxmartijn@gmail.com>
Co-authored-by: Claude Code <noreply@anthropic.com>
Co-authored-by: Alex Reinking <areinking@adobe.com>
Co-authored-by: halide-ci[bot] <266445882+halide-ci[bot]@users.noreply.github.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants